Softmax优化:从Naive到两遍融合
测试平台:RTX 5080 (Blackwell, sm_120),M=N=8192,理论带宽960 GB/s
前言:softmax是attention的核心算子,也是典型memory-bound算子。深度学习框架中softmax算子几乎是必不可少的组件,在实现时一般也是GOU kernel function形式。
| 版本 | 核心优化 | 配置 | 带宽(GB/s) | 相对v1 |
| --- | ---------------------- | ------------ | ---------- | ------- |
| v1 | naive:一个线程一行 | 1×256 | ~150 | 基准 |
| v2 | shared memory树形归约 | 256×1 | ~820 | ~5.5× |
| v3 | warp shuffle两级归约 | 256×1 | ~810 | 持平 |
| v4 | float4向量化加载 | 256×1 | ~820 | 略升 |
| v5 | online两遍融合 | 256×1 | ~800 | 持平 |
v2起各版本都在~85%峰值附近,真正的瓶颈可能是DRAM带宽,无法再突破DRAM上限。但5080显卡无dram读取计数器,无法获取真实带宽利用率。
1 Safe Softmax与三遍扫描
y_i = exp(x_i - max) / Σ exp(x_j - max)
三遍扫描:
- 第一遍求行最大值
max - 第二遍求指数和
sum - 第三遍归一化写出
y
2 v1版本:朴素实现
一个线程负责一行,串行扫三遍:
__global__ void softmax_v1_naive(float *input, float *output, int M, int N) {
int row = blockIdx.y * blockDim.y + threadIdx.y;
if (row >= M) return;
float *x = input + row * N;
float *y = output + row * N;
// 求最大值
float max = -INFINITY;
for (int k = 0; k < N; k++) max = fmax(max, x[k]);
// 指数求和
float sum = 0.0f;
for (int k = 0; k < N; k++) sum += expf(x[k] - max);
// 归一化
float inv_sum = 1.0f / sum;
for (int k = 0; k < N; k++) y[k] = expf(x[k] - max) * inv_sum;
}
结果~150 GB/s,只有峰值的15%。
分析:
- 并行度太低:一个线程要串行扫完整行8192个元素,8192次循环。
- 非合并访存:相邻线程(threadIdx.y相邻)处理的是不同行,它们访问的地址相差N(32KB),一个warp的32次load落在32个不同cache line上,非合并访存。
3 v2版本:一个block一行 + shared memory归约
让一个block(256线程)协作处理一行,每个线程只负责一行的一部分,再在block内做归约。归约用shared memory做树形规约:
__global__ void softmax_v2_sharedreduce(float *input, float *output, int M, int N) {
int row = blockIdx.y; // 一个block处理一行
int tx = threadIdx.x;
if (row >= M) return;
float *x = input + row * N;
float *y = output + row * N;
extern __shared__ float smem[];
// Pass 1:并行求最大值
float max_val = -INFINITY;
for (int i = tx; i < N; i += blockDim.x) max_val = fmaxf(max_val, x[i]);
smem[tx] = max_val;
__syncthreads();
for (int s = blockDim.x / 2; s > 0; s >>= 1) {
if (tx < s) smem[tx] = fmaxf(smem[tx], smem[tx + s]);
__syncthreads();
}
max_val = smem[0];
__syncthreads();
// Pass 2 / Pass 3:求sum、归一化写出,结构同上
}
结果跳到~820 GB/s,是v1的5.5倍。
分析:
- 每个线程用
for (i = tx; i < N; i += blockDim.x)跨步访问,256个线程把一行铺满,访存变成合并的。 - 归约用树形
for (s = blockDim.x/2; s > 0; s >>= 1),256 → 1 只需要8步,每步一次__syncthreads()。
如果sum归约的循环条件写成 for (int i = blockDim.x/2; i < N; i >>= 1),一方面 i < N永远成立会死循环,另一方面 smem[tx + i]会越界。正确写法是上面的 s >>= 1且 s > 0且 if (tx < s)保护。
4 v3版本:warp shuffle减少同步
v2的树形归约每一步都要 __syncthreads()(整个block同步8次),而warp内的32个线程是天然同步的。warp内用 __shfl_down_sync 归约(寄存器级、免同步),warp之间再用一小块shared memory归约,同步次数从8次降到2次。
// warp内归约:5次shuffle,无同步
__device__ float warpReduceMax(float val) {
for (int offset = 16; offset > 0; offset >>= 1)
val = fmaxf(val, __shfl_down_sync(0xffffffff, val, offset));
return val;
}
// block归约:warp内shuffle + warp间smem
__device__ float blockReduceMax(float val) {
__shared__ float warp_max[32];
int lane = threadIdx.x & 31;
int wid = threadIdx.x >> 5;
val = warpReduceMax(val);
if (lane == 0) warp_max[wid] = val;
__syncthreads(); // 仅1次同步收集各warp
int nwarps = blockDim.x / 32;
val = (lane < nwarps) ? warp_max[lane] : -INFINITY;
if (wid == 0) val = warpReduceMax(val); // 第0个warp收尾
__shared__ float result;
if (threadIdx.x == 0) result = val;
__syncthreads();
return result;
}
结果~810 GB/s,和v2持平。
5 v4版本:float4向量化加载
既然瓶颈在访存,就把load/store指令数降下来:一次读4个float(128-bit的 float4),三遍循环全部向量化。
const float4* x4 = reinterpret_cast<const float4*>(input + row * N);
float4* y4 = reinterpret_cast<float4*>(output + row * N);
int N4 = N / 4; // 计数单位从N变成N/4
// Pass 1:一次读4个,取4分量最大值
float max_val = -INFINITY;
for (int i = tx; i < N4; i += blockDim.x) {
float4 v = x4[i];
max_val = fmaxf(max_val, fmaxf(fmaxf(v.x, v.y), fmaxf(v.z, v.w)));
}
// Pass 2 / Pass 3同理,4个分量分别exp后累加 / 构造float4一次写出
分析:
- load/store指令从4条
LDG.32变成1条LDG.128,指令数大幅下降。 - 4个连续线程访问4个连续float4,正好凑满128字节cache line,合并度更好。
结果~820 GB/s,相对v3略升。收益有限,因为DRAM已经接近满负荷,指令数减少只能缓解,无法突破带宽上限。
6 v5版本:两遍融合(online softmax)
三遍扫描读input三次,能不能把max和sum合并到一次读取里,只读两次?这就是online softmax的思路。
问题在于:sum依赖最终的max,而max是扫到最后才知道的。解决办法是维护一个running (m, s),遇到更大的值就把旧的sum按指数缩放:
遇到新元素x,设 m_new = max(m, x):
s = s * exp(m - m_new) + exp(x - m_new)
m = m_new
当 x <= m 时 exp(m - m_new) = 1,等价于 s += exp(x - m);当 x > m 时,旧sum乘以 exp(m - x) 缩放到新基准。
block内每个线程先算自己的局部(m, s),归约时先归约出全局max,再把每个线程的s缩放到全局max后求和:
// Pass 1:一遍读取,在线更新(m, s)
float m = -INFINITY, s = 0.0f;
for (int i = tx; i < N4; i += blockDim.x) {
float4 v = x4[i];
float mn;
mn = fmaxf(m, v.x); s = s * expf(m - mn) + expf(v.x - mn); m = mn;
mn = fmaxf(m, v.y); s = s * expf(m - mn) + expf(v.y - mn); m = mn;
mn = fmaxf(m, v.z); s = s * expf(m - mn) + expf(v.z - mn); m = mn;
mn = fmaxf(m, v.w); s = s * expf(m - mn) + expf(v.w - mn); m = mn;
}
float global_m = blockReduceMax(m);
s = s * expf(m - global_m); // 缩放到全局max
float global_s = blockReduceSum(s);
// Pass 2:一遍读取,归一化写出
结果~800 GB/s,没有比v4快。一开始我以为两遍融合能省掉一次DRAM读,应该快1/3,结果没变,于是用ncu查了真实数据。
7 用ncu验证:瓶颈是DRAM带宽
ncu数据(8192×8192,1 sector = 32字节,读一遍 = 256 MiB = 8,388,608 sectors):
v4(三遍扫描):
gpu__time_duration.sum us 622.88
lts__t_bytes.sum Gbyte 1.08
lts__t_sectors_op_read.sum sector 25,331,031 ≈ 3次
lts__t_sectors_op_write.sum sector 8,388,623 ≈ 1次
dram__bytes_read.sum (!) n/a
dram__bytes_write.sum (!) n/a
v5(两遍融合):
gpu__time_duration.sum us 631.20
lts__t_bytes.sum Mbyte 820.83
lts__t_sectors_op_read.sum sector 17,260,927 ≈ 2次
lts__t_sectors_op_write.sum sector 8,388,620 ≈ 1次
dram__bytes_read.sum (!) n/a
dram__bytes_write.sum (!) n/a
gpu__time_duration.sum:kernel执行时长(微秒)。lts__t_sectors_op_read/write.sum:到达L2缓存的读/写请求量,单位sector(32字节)。lts即L2 slice。dram__bytes_read/write.sum:DRAM流量,GeForce消费卡不开放这个计数器,所以是n/a。
对比:
- L2读流量:v4是v5的1.47倍(25.3M vs 17.3M sectors),证明v5确实省掉了一次读。
- 但两者时间几乎一样(622 vs 631 us)。
结论:省掉的那次读命中了L2,根本没打到DRAM。一行32KB,同一个block里三个pass间隔极短,数据一直躺在L2里。所以三遍扫描的DRAM真实流量早就已经是下限「读1次 + 写1次 = 512 MiB」,两遍融合省的是L2带宽(L2不紧张),不省DRAM,因此不加速。
用时间反推真实DRAM带宽:
512 MiB / 622.88 us ≈ 862 GB/s (约89%峰值)
8 总结与相关参考
参考博客:
特别鸣谢:Deepseek-v4-Pro、Trae Agent
